Skip to content

feat(jit): forward sub-part geometry env vars to the JIT - #1

Open
whn09 wants to merge 1 commit into
amazon-contributing:mainfrom
whn09:jit-forward-subpart-env
Open

feat(jit): forward sub-part geometry env vars to the JIT#1
whn09 wants to merge 1 commit into
amazon-contributing:mainfrom
whn09:jit-forward-subpart-env

Conversation

@whn09

@whn09 whn09 commented Aug 21, 2026

Copy link
Copy Markdown

Re-post of Xuan-1998/DeepEP#42 against this repo, rebased onto main (ec623f3) and re-measured on that tree — the old PR's numbers were taken on the pre-refactor tree, so they are not quoted here.

What

hybrid_dispatch_unordered.cuh gates its sub-part geometry behind #ifndef:

macro default
EP_NUM_SUB_PARTS 2
EP_MIN_SUB_TOKENS 1
EP_SM100_MIN_SUB_TOKENS 15

Nothing in the tree ever sets those macros, so the only way to try a different split today is to edit the header and reinstall. This forwards the three names as JIT -D flags, following the EP_NUM_TOPK_IDX_BITS block immediately above it (and that block's EP_JIT_EXTRA_FLAGS TODO).

Unset ⇒ byte-identical behaviour.

Why it is safe

  • All three are device-only. No host translation unit reads them, so a JIT-only define cannot desync host-side sizing from device-side sizing. (This is the reason the patch stops at these three and does not forward, say, kNumMaxTokensPerRank.)
  • flags is part of kernel_signature, so changing the env re-JITs instead of silently serving a cached cubin compiled with a different geometry.

Verified end-to-end

EP_JIT_PRINT_COMPILER_COMMAND=1 EP_NUM_SUB_PARTS=1 puts -DEP_NUM_SUB_PARTS=1 on the nvcc line, and the resulting cubin lands in a distinct JIT cache entry.

Measured effect of the knob it exposes

tests/elastic/test_ep.py, 2 × p5en.48xlarge (8×H200 + 16 EFA each), EP8×2 = 16 ranks, --hidden=7168 --num-topk=8 --num-experts=256 --num-sms=12 --allow-hybrid-mode=1 --prefer-overlap-with-compute=0 --test-first-only. 3 reps, variants interleaved within each rep, each variant on its own EP_JIT_CACHE_DIR. Mean over all 16 ranks then over reps; ± is stdev across reps. All 48 rounds exited 0.

EP_NUM_SUB_PARTS=1 on its own:

op baseline (unset) EP_NUM_SUB_PARTS=1
prefill (8192 tok) dispatch 1665.1 ± 12.4 µs 1631.9 ± 2.8 µs −2.0%
prefill cached dispatch 1662.8 ± 8.7 µs 1635.7 ± 1.2 µs −1.6%
prefill combine 3560.8 ± 9.2 µs 3545.8 ± 5.0 µs −0.4%
decode (128 tok) dispatch 367.0 ± 12.1 µs 373.7 ± 3.6 µs +1.8%
decode combine 178.1 ± 5.0 µs 178.5 ± 0.3 µs +0.2%

So on its own the knob is roughly a wash — a small prefill win, a small decode regression. Its value is that it composes: stacked with #2 it takes decode dispatch from 367.0 µs to 166.1 ± 0.4 µs (−54.7%), where that PR alone reaches 239.8 µs (−34.7%). That is the case this PR is really enabling — being able to find such a combination without a rebuild.

This PR is pure plumbing and changes no default, so it is worth taking independently of whether #2 is accepted.

`hybrid_dispatch_unordered.cuh` gates the sub-part geometry behind `#ifndef`
(`EP_NUM_SUB_PARTS` 2, `EP_MIN_SUB_TOKENS` 1, `EP_SM100_MIN_SUB_TOKENS` 15), but
nothing in the tree sets those macros, so the only way to try a different split
is to edit the header and reinstall. Forward the three names as JIT `-D` flags,
following the `EP_NUM_TOPK_IDX_BITS` block immediately above (and its
`EP_JIT_EXTRA_FLAGS` TODO).

All three are device-only -- no host translation unit reads them -- so a
JIT-only define cannot desync host and device sizing. `flags` is part of
`kernel_signature`, so changing the env re-JITs instead of serving a cached
cubin. Unset => no behaviour change.
@whn09
whn09 force-pushed the jit-forward-subpart-env branch from efe26c2 to 5118c2e Compare August 24, 2026 06:26
@whn09

whn09 commented Aug 24, 2026

Copy link
Copy Markdown
Author

Rebased onto main @ 02efc268; force-pushed. The change itself is unmodifiedgit patch-id --stable on the rebased commit matches the original.

The old branch showed a 3,000-line diff over 20 files, which was an artefact, not this change. It was cut from main @ ec623f3 (committed 08-21 00:05 UTC); main was then rewritten and force-pushed (cc55cce, committed 08-21 23:51 UTC, same author date). The four base commits exist in both histories under different SHAs, so ec623f3 became unreachable from main and the merge-base fell back six commits to 01dc3aa, attributing the whole unordered-kernel feature to this PR. Both files this PR touches are byte-identical between ec623f3 and main, so the replay was conflict-free.

Measurement caveat: the numbers in the description were taken on ec623f3, not on the rebased tree. ec623f3 → main is +124/−20 across six files, including csrc/kernels/backend/nccl.cu and csrc/kernels/elastic/combine.hpp, so I would not claim they carry over unchanged. The p5en pair used for them is no longer available to me; happy to re-run if you would like the numbers refreshed on main before merging.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants